热门标签 | HotTags
当前位置:  开发笔记 > 编程语言 > 正文

也就是|小窗_卷积的特征提取与参数计算

篇首语:本文由编程笔记#小编为大家整理,主要介绍了卷积的特征提取与参数计算相关的知识,希望对你有一定的参考价值。Dense和Conv2D根本区别在于,Den

篇首语:本文由编程笔记#小编为大家整理,主要介绍了卷积的特征提取与参数计算相关的知识,希望对你有一定的参考价值。


Dense和Conv2D根本区别在于,Dense层从输入空间中学到的是全局模式,比如对于MNIST数字来说,全局模式就是涉及所有像素的模式。而Conv2D学到的是局部模式(local pattern),同样以MNIST为例,Conv2D学到的是在输入图像的小窗口中发现的模式(pattern)。

这个重要特性使卷积神经网络具有两个有趣的性质:

1. 卷积神经网络学到的模式具有平移不变性。卷积神经网络在图像右下角学到某个模式后,它可以在任何地方识别这个模式,比如左上角。但是对于Dense网络来说,如果模式出现在新的位置那么必须重新学习这个模式。这使得卷积神经网络在处理图像时可以更高效的利用数据。

2. 卷积神经网络可以学到模式的空间层次结构。如下图所示,第一个卷积层将学习较小的局部模式(比如边缘),第二个卷积层将学习由第一个层的特征们组成的更大的模式,以此类推,使得卷积神经网络可以学习越来越复杂,越来越抽象的视觉概念。

对于MINST,第一个卷积层接收一个大小为(28,28,1)的输入特征图。为了通过卷积从该输入特征图中提取不同的局部模式,我们要设计不同的卷积核。每个卷积核的大小宽高多为3*3或者5*5,卷积核的深度与输入特征图的深度一致,卷积核的个数与这一层需要获得的局部模式数量一致。每一个卷积核也成为一个滤波器,通过滤波器的过滤(filter),就学到了一个局部模式(特征)。例如,我设计的第一个卷积层希望从这张图片中获取32种局部模式,那么我就要设计32个(3*3*1)的卷积核,经过第一层过滤后,会形成一个(26*26*32)的输出特征图。每个过滤器与输入特征图进行卷积运算会得到一个(26*26)的响应图(response map),32个过滤器就会得到(26*26*32)这样输出特征图。

所以,特征图深度方向上的每一个维度都是一个特征(过滤器),而每一个维度上的2D张量是该维度过滤器对输入的响应所形成的二维空间图(map)。

 现在,对第一次卷积层得到的(26*26*32)特征图进行一次最大值池化(MaxPooling),经过池化的特征数量变为(13*13*32)。随后设计第二个卷积层,对于第二个卷积层而言输入的特征图就是(13,13,32),那么第二个卷积层的卷积核深度也应该是32,也就是(3,3,32),比如第二个卷积层我想提取64个特征,那么就意味着这一层经过运算后,输出的特征图为(11,11,64),随后是针对第二层卷积输出的最大值池化,经过池化的特征数量变为(5,5,64),然后可以设计第三个卷积层,那么根据输入特征图的深度,第三个卷积层的卷积核深度为64,也就是(3,3,64),比如第三个卷积层我也想提取64个特征,那么就意味着这一层经过运算后,输出的特征图为(3,3,64)。经过第三次卷积后,将输出内容可以与Dense层连接,然后分类输出。当然在于Dense层连接前,需要将3维数据“抻平”,变为一维数据才能输入到Dense层。最终我们以代码的形式来设计一下文中所述的神经网络:

from tensorflow import keras
from tensorflow.keras import models
from tensorflow.keras import layers
from tensorflow.keras import datasets
from tensorflow.keras import utils
mymodel = models.Sequential()
mymodel.add(layers.Conv2D(32,(3,3),activation='relu',input_shape=(28,28,1)))
mymodel.add(layers.MaxPooling2D(2,2))
mymodel.add(layers.Conv2D(64,(3,3),activation='relu'))
mymodel.add(layers.MaxPooling2D(2,2))
mymodel.add(layers.Conv2D(64,(3,3),activation='relu'))
mymodel.add(layers.Flatten())
mymodel.add(layers.Dense(64,activation='relu'))
mymodel.add(layers.Dense(10,activation='softmax'))
mymodel.summary()

 模型的summary函数会返回每个层的参数量:



_________________________________________________________________
 Layer (type)                Output Shape              Param #
=================================================================
 conv2d (Conv2D)             (None, 26, 26, 32)        320


 max_pooling2d (MaxPooling2D  (None, 13, 13, 32)       0
 )


 conv2d_1 (Conv2D)           (None, 11, 11, 64)        18496


 max_pooling2d_1 (MaxPooling  (None, 5, 5, 64)         0
 2D)


 conv2d_2 (Conv2D)           (None, 3, 3, 64)          36928


 flatten (Flatten)           (None, 576)               0


 dense (Dense)               (None, 64)                36928


 dense_1 (Dense)             (None, 10)                650


如果能理解卷积网络层中核的形状和个数,那么每一层的参数就不难理解了。



 conv2d (Conv2D)             (None, 26, 26, 32)        320


第一个卷积层的卷积核大小为(3,3,1),共32个,所以w = 3*3*1*32 = 288,有w就会有偏置b,一维向量共32个元素。所以一共有320个参数。 



 conv2d_1 (Conv2D)           (None, 11, 11, 64)        18496


第二个卷积层的卷积核大小为(3,3,32),共64个,所以w = 3 * 3 *32 * 64 =  18432,还有偏置b中的64个变量,所以一共有18496个参数。



 conv2d_2 (Conv2D)           (None, 3, 3, 64)          36928


第三个卷积层的卷积核大小为(3,3,64),共64个,所以w = 3 * 3 *64 * 64 =  36864,还有偏置b中的64个变量,所以一共有36928个参数。



 dense (Dense)               (None, 64)                36928


输入层维度为576(把(3,3,64)给抻平),输出层维度是64,所以w是 576*64 = 36864,再加上偏置向量中的64个元素,所以一共有36928个参数。



 dense_1 (Dense)             (None, 10)                650


最后一个Dense层,输入为64,输出是10,所以w是64*10 = 640,最后再加上10个偏置,所以一共有650个参数。

整个神经网络的参数总数为93322个。

(train_img,train_labels),(test_img,test_labels) = datasets.mnist.load_data()
train_img = train_img.reshape((60000,28,28,1))
train_img = train_img.astype('float32')/255
test_img = test_img.reshape((10000,28,28,1))
test_img = test_img.astype('float32')/255
train_labels = utils.to_categorical(train_labels)
test_labels = utils.to_categorical(test_labels)
mymodel.compile(optimizer='rmsprop',loss='categorical_crossentropy',metrics=['accuracy'])
mymodel.fit(train_img,train_labels,epochs=5,batch_size=64)
myloss,myaccuracy =mymodel.evaluate(test_img,test_labels)
print(f'模型测试准确率:myaccuracy')

最终在测试集上得到的准确率为0.9912999868392944。

相较于单纯使用Dense进行组网准确率还是有明显上升的(纯Dense组网的准确率在97%以上)。


推荐阅读
  • 深入浅出TensorFlow数据读写机制
    本文详细介绍TensorFlow中的数据读写操作,包括TFRecord文件的创建与读取,以及数据集(dataset)的相关概念和使用方法。 ... [详细]
  • 尽管深度学习带来了广泛的应用前景,其训练通常需要强大的计算资源。然而,并非所有开发者都能负担得起高性能服务器或专用硬件。本文探讨了如何在有限的硬件条件下(如ARM CPU)高效运行深度神经网络,特别是通过选择合适的工具和框架来加速模型推理。 ... [详细]
  • Python实现照片磨皮效果
    本文介绍如何使用Python和OpenCV库来实现照片的磨皮效果,使图片更加平滑并提升整体美感。 ... [详细]
  • 卷积神经网络(CNN)基础理论与架构解析
    本文介绍了卷积神经网络(CNN)的基本概念、常见结构及其各层的功能。重点讨论了LeNet-5、AlexNet、ZFNet、VGGNet和ResNet等经典模型,并详细解释了输入层、卷积层、激活层、池化层和全连接层的工作原理及优化方法。 ... [详细]
  • 基于2-channelnetwork的图片相似度判别一、相关理论本篇博文主要讲解2015年CVPR的一篇关于图像相似度计算的文章:《LearningtoCompar ... [详细]
  • 在Ubuntu 16.04中使用Anaconda安装TensorFlow
    本文详细介绍了如何在Ubuntu 16.04系统上通过Anaconda环境管理工具安装TensorFlow。首先,需要下载并安装Anaconda,然后配置环境变量以确保系统能够识别Anaconda命令。接着,创建一个特定的Python环境用于安装TensorFlow,并通过指定的镜像源加速安装过程。最后,通过一个简单的线性回归示例验证TensorFlow的安装是否成功。 ... [详细]
  • 毕业设计:基于机器学习与深度学习的垃圾邮件(短信)分类算法实现
    本文详细介绍了如何使用机器学习和深度学习技术对垃圾邮件和短信进行分类。内容涵盖从数据集介绍、预处理、特征提取到模型训练与评估的完整流程,并提供了具体的代码示例和实验结果。 ... [详细]
  • 推荐几款高效测量图片像素的工具
    本文介绍了几款适用于Web前端开发的工具,这些工具可以帮助用户在图片上绘制线条并精确测量其像素长度。对于需要进行图像处理或设计工作的开发者来说非常实用。 ... [详细]
  • 本文介绍了如何利用TensorFlow框架构建一个简单的非线性回归模型。通过生成200个随机数据点进行训练,模型能够学习并预测这些数据点的非线性关系。 ... [详细]
  • 本文探讨了亚马逊Go如何通过技术创新推动零售业的发展,以及面临的市场和隐私挑战。同时,介绍了亚马逊最新的‘刷手支付’技术及其潜在影响。 ... [详细]
  • 由二叉树到贪心算法
    二叉树很重要树是数据结构中的重中之重,尤其以各类二叉树为学习的难点。单就面试而言,在 ... [详细]
  • 使用OpenCV和Python 4.2提升模糊图像清晰度
    本文介绍如何利用OpenCV库在Python中处理图像,特别是通过不同类型的滤波器来改善模糊图像的质量。我们将探讨均值、中值和自定义滤波器的应用,并展示代码示例。 ... [详细]
  • 2017年人工智能领域的十大里程碑事件回顾
    随着2018年的临近,我们一同回顾过去一年中人工智能领域的重要进展。这一年,无论是政策层面的支持,还是技术上的突破,都显示了人工智能发展的迅猛势头。以下是精选的2017年人工智能领域最具影响力的事件。 ... [详细]
  • 本文档旨在帮助开发者回顾游戏开发中的人工智能技术,涵盖移动算法、群聚行为、路径规划、脚本AI、有限状态机、模糊逻辑、规则式AI、概率论与贝叶斯技术、神经网络及遗传算法等内容。 ... [详细]
  • 利用Java与Tesseract-OCR实现数字识别
    本文深入探讨了如何利用Java语言结合Tesseract-OCR技术来实现图像中的数字识别功能,旨在为开发者提供详细的指导和实践案例。 ... [详细]
author-avatar
mobiledu2502856013
这个家伙很懒,什么也没留下!
PHP1.CN | 中国最专业的PHP中文社区 | DevBox开发工具箱 | json解析格式化 |PHP资讯 | PHP教程 | 数据库技术 | 服务器技术 | 前端开发技术 | PHP框架 | 开发工具 | 在线工具
Copyright © 1998 - 2020 PHP1.CN. All Rights Reserved | 京公网安备 11010802041100号 | 京ICP备19059560号-4 | PHP1.CN 第一PHP社区 版权所有